KingOfEgg
首页项目归档照片墙音乐灵境说说杂谈友链关于
封面

Softmax优化(二) Online Softmax从V0到Grid Stride

写作时间:2026-08-19 12:00:00
# CUDA
# Softmax
# Online Softmax
# FlashAttention
# GPU优化

Softmax优化(二) Online Softmax从V0到Grid Stride

测试平台:RTX 5080 (Blackwell, sm_120),M=N=4096,代码在 cuda_code/softmax/online-softmax

前言:笔记一《Softmax优化(一)从Naive到两遍融合》走完了Safe Softmax三遍扫描到两遍融合的路径,结论是两遍融合在RTX 5080上没有提速——省掉的那次读命中了L2而非DRAM。这篇接着AIInfraGuide的教程,从算法推导出发,把Online Softmax的CUDA Kernel从V0单线程一路实现到V4 Grid Stride,重点放在online递推的并行化与寄存器缓存这两个Safe Softmax没有的新问题。

版本核心优化配置时间(ms)相对V0V0单线程一行1×2561.23基准V1block并行+warp shuffle合并规约256×10.1677.4×V2float4向量化256×10.166持平V3寄存器缓存one-pass256×10.165持平V4grid stride256×10.173持平

V1之后各版本都在~0.17ms附近打平,与笔记一的结论一致:瓶颈是DRAM带宽,V1已经把访存压到位,后续优化省的都是L2命中或寄存器,不省DRAM。

1 Online Softmax的递推与并行化

Safe Softmax要三遍扫描(max→sum→归一化),Online Softmax把max和sum合并到一遍,读次数从3降到2。递推维护running (m,d):

遇到新元素x,设 m_new = max(m, x):
    d = d * exp(m - m_new) + exp(x - m_new)
    m = m_new

当x<=m时exp(m-m_new)=1,退化为普通累加;当x>m时旧sum乘以exp(m-x)缩放到新基准。

递推本身是串行的,但可以并行化,关键是两个局部结果(m1,d1)、(m2,d2)的合并操作满足结合律:

m_merged = max(m1, m2)
d_merged = d1 * exp(m1 - m_merged) + d2 * exp(m2 - m_merged)

所以可以先让每个线程算自己那段的局部(m,d),再用树形规约合并。这正是V1的基础。

2 V0:单线程Online Softmax

一个线程串行处理一行,直接按上面的递推扫一遍,再归一化:

__global__ void online_softmax_v0_naive(const float* input, float* output, int M, int N) {
    int row = blockIdx.y * blockDim.y + threadIdx.y;
    if (row >= M) return;
    const float* x = input + row * N;
    float* y = output + row * N;
    // 第1遍:online同时更新(m,d)
    float m = -INFINITY, d = 0.0f;
    for (int i = 0; i < N; i++) {
        float m_new = fmaxf(m, x[i]);
        d = d * expf(m - m_new) + expf(x[i] - m_new);
        m = m_new;
    }
    // 第2遍:归一化
    float inv_d = 1.0f / d;
    for (int i = 0; i < N; i++) y[i] = expf(x[i] - m) * inv_d;
}

结果~1.23ms,约163 GB/s。和Safe Softmax的V0一样,单线程独享一行,并行度为零、非合并访存,只作正确性基线。

3 V1:Block并行+warp shuffle合并规约

一个block处理一行,每个线程先串行算自己那段的局部(m,d),再做两级归约:warp内用__shfl_down_sync合并,warp间经shared memory交换,warp0收尾。

__device__ void warpReduceOnline(float& m, float& d) {
    for (int offset = 16; offset > 0; offset >>= 1) {
        float m2 = __shfl_down_sync(0xffffffff, m, offset);
        float d2 = __shfl_down_sync(0xffffffff, d, offset);
        float m_new = fmaxf(m, m2);
        d = d * expf(m - m_new) + d2 * expf(m2 - m_new);
        m = m_new;
    }
}

__global__ void online_softmax_v1(...) {
    // 每线程局部(m,d)
    float local_m = -INFINITY, local_d = 0.0f;
    for (int i = tid; i < N; i += blockDim.x) { ... }  // online递推
    // warp内合并 + warp间smem交换 + warp0收尾
    warpReduceOnline(local_m, local_d);
    ...
    // 第2遍:归一化
}

结果~0.167ms,约1208 GB/s,比V0快7.4倍。

4 V2:float4向量化

算法不变,只优化访存:把float重解释成float4,一次LDG.128读4个、一次STG.128写4个。但online递推仍是逐元素更新(m,d),展开4个分量依次做。

const float4* x4 = reinterpret_cast<const float4*>(x);
float4* y4 = reinterpret_cast<float4*>(y);
int N4 = N / 4;
// 第1遍:一次读4个,逐个更新(m,d)
float4 v = x4[i];
float mn;
mn = fmaxf(local_m, v.x); local_d = local_d * expf(local_m - mn) + expf(v.x - mn); local_m = mn;
// v.y / v.z / v.w同理

结果~0.166ms,与V1持平。访存指令数÷4、循环迭代÷4,但DRAM已经接近满负荷,指令减少只能缓解不能突破带宽上限。

5 V3:寄存器缓存one-pass(编译期展开的坑)

V1/V2都读2遍input,V3想把第1遍读到的元素缓存到寄存器,第2遍直接读寄存器,把全局访存压到1读+1写。

第一版用运行时下标reg_cache[count],结果ptxas显示缓存被放到了stack(local memory)而不是寄存器:

// 运行时下标:128 bytes stack frame, 25 registers
float reg_cache[MAX_ELEMS];
for (int i = tid; i < N; i += blockDim.x) {
    reg_cache[count++] = x[i];   // count是运行时变量
}

要真正留在寄存器,下标必须是编译期常量。把每线程元素数ELEMS做成模板参数,循环用#pragma unroll完全展开:

template <int ELEMS>
__global__ void online_softmax_v3_regcache(...) {
    float reg_cache[ELEMS];
    #pragma unroll
    for (int k = 0; k < ELEMS; k++) {
        int i = tid + k * blockDim.x;
        float xi = (i < N) ? x[i] : -INFINITY;
        reg_cache[k] = xi;   // k是编译期常量→寄存器
        ...
    }
}

ptxas对比:展开后ELEMS=16时0 bytes stack frame、56 registers,缓存真正进了寄存器(ELEMS=32是88,ELEMS=64是244)。

结果~0.165ms,仍然与V2持平。原因和笔记一一致:V2省不掉的那次读本来就命中L2,V3从L2再压到寄存器收益更小,反而寄存器压力从34涨到56,占用率下降抵消了省读收益。

6 V4:Grid Stride

前面都是一block一行(grid=M),M很大时会启动海量block、调度开销线性增加。V4用固定block数循环覆盖所有行,算法与V2相同,区别只在block与行的映射从1:1变成1:多。

for (int row = blockIdx.x; row < M; row += gridDim.x) {
    // 一整行:online递推 + warp shuffle合并 + 归一化,与V2相同
}

共享内存数组移到循环外跨行复用,grid=min(M, sm_count×4)。结果~0.173ms,与V2持平——M=4096时DRAM已打满,调度不是瓶颈,grid stride的收益只在M极大时体现。

7 为什么V1之后全打平

整条路径下来,V1之后全部卡在~0.17ms。结合笔记一的ncu结论:这张卡的softmax是memory-bound,V1已经做到「读2次+写1次」且额外读命中L2,真实DRAM流量早就到了下限「读1次+写1次」,后续的float4、寄存器缓存、grid stride省的都是L2带宽或调度开销,不省DRAM,所以不加速。

8 总结与FlashAttention的联系

单看「写个softmax kernel」,优化招式(合并访存、向量化、规约、寄存器分块、grid stride)都是GEMM那套工具库的复用。Online Softmax真正的新东西是算法层面:running (m,d)递推 + 合并结合律,这才让并行规约成为可能。它的价值不在「省读」,而在FlashAttention里不物化N×N分数矩阵——只能靠online softmax边算边归一化,跨KV block重缩放。

参考:

AIInfraGuide:CUDA Online Softmax实现

Softmax优化(一)从Naive到两遍融合

avatar

KingOfEgg

otaku change the world

RECOMMENDED

GEMM 优化(一) Cuda Core计算

2026-08-11 12:00:00

CUBLAS GEMM 函数用法

2026-08-12 12:00:00

GEMM优化(三)大矩阵适配和极致优化

2026-09-03 12:00:00

Table of Contents